Attention Architecture Evolution

导言

Attention 的发展并不是从 Full Attention 排成一条单向淘汰链。更准确的结构是三条并行路线:共享或压缩 KV cache、把历史压入固定状态、对历史 token 做稀疏选择。MQA、GQA 与 MLA 仍然读取全部历史;GDN 与 KDA 改写了记忆算法;DSA 则在 MLA 前增加轻量索引器,只让主注意力读取 top-k。

本文以首个公开论文或官方发布日为时间点,并把每种结构落到版本固定的开源代码:Q/K/V 从哪里产生、什么对象进入 cache、score 如何形成、复杂度到底被搬到了哪里。

结论

先给出最重要的判断:

  1. Full/MHA、MQA、GQA、MLA 是同一 dense softmax 家族中的不同 KV 组织方式。后三者主要减少 KV cache 或带宽,并不自动把全部 token 两两交互变成线性复杂度。
  2. GDN 与 KDA 属于递归状态路线。历史不再以逐 token K/V 保存,而是写入固定大小矩阵状态;KDA 把 GDN 的 head-wise 标量遗忘细化为 key-channel-wise 对角遗忘。
  3. DSA 属于 token 稀疏路线。Lightning Indexer 仍扫描历史并产生 top-k,主 MLA 才只读取选中位置,所以“主注意力稀疏”不等于整个模块没有历史扫描成本。
  4. 公开时间并不等于概念顺序。DSA 于 2025 年 9 月公开,KDA 于 2025 年 10 月公开;本文按真实时间线排列,而不是按缩写的相似度排列。
![Attention 的三条演进路线](https://pic.shaojiemike.top/shaojiemike/2026/07/126eff6e8371c0195d2790b2bf974eec.png){ width=94% }
小黑认知示意图:Attention 优化的核心不是统一“少算一点”,而是在共享、压缩、固定状态与稀疏选择之间重新组织历史记忆。

时间线

首次公开 结构 首个一手来源 主要变化 历史表示
2017-06-12 Full Attention / MHA Attention Is All You Need[^mha-paper] 每个 Query head 拥有独立 K/V,读取全部历史 完整 K/V cache
2019-11-06 MQA Fast Transformer Decoding: One Write-Head is All You Need[^mqa-paper] 多个 Q head 共享一组 K/V 单组 K/V cache
2023-05-22 GQA GQA: Training Generalized Multi-Query Transformer Models[^gqa-paper] Q head 分组共享 K/V 多组 K/V cache
2024-05-07 MLA DeepSeek-V2[^mla-paper] K/V 联合低秩压缩,解耦 RoPE cKV + kR
2024-12-09 GDN Gated Delta Networks[^gdn-paper] head-wise 衰减的 delta-rule 固定状态 S ∈ R^(K×V)
2025-09-29 DSA DeepSeek-V3.2-Exp[^dsa-release] Indexer 选 top-k,主 MLA 稀疏读取 MLA cache + Indexer K
2025-10-30 KDA Kimi Linear[^kda-paper] key-channel-wise 对角衰减的 delta-rule S ∈ R^(K×V)

Full Attention 与 MHA

Full Attention 描述 token 连通性:每个位置可读取全部允许的历史位置;MHA 描述 head 组织:每个 Query head 拥有独立 K/V head。自回归解码器通常同时使用 causal Full Attention 和 MHA,因此工程讨论里经常连写,但两者不是完全同义词。

三条演进路线

定义:序列长度为 T,Query head 数为 Hq,KV head 数为 Hkv,head 维度为 d,MLA 的 KV latent 维度为 r,DSA 选出的历史位置数为 k

结构 Token 连通性 T 增长的主要缓存 单步 decode 主计算 优化焦点
Full/MHA Dense 2 × T × Hq × d O(T × Hq × d) 基线
MQA Dense 2 × T × d O(T × Hq × d) KV cache 与带宽
GQA Dense 2 × T × Hkv × d O(T × Hq × d) 质量与 KV 成本折中
MLA Dense T × (r + d_rope) 仍扫描全部历史 latent 低秩 KV cache
GDN Recurrent T 无关的 H × K × V 状态 O(H × K × V) 固定状态
KDA Recurrent T 无关的 Hv × K × V 状态 O(Hv × K × V) 更细粒度遗忘
DSA 主注意力 top-k MLA cache + Indexer K Indexer 扫描 T,主 MLA 读取 k 长上下文主计算稀疏化

这张表解释了一个常见误区:KV cache 更小、token 连通性更稀疏、算法复杂度更低是三个不同命题。MQA/GQA 改第一个,DSA 改第二个,GDN/KDA 改第三个所依赖的历史表示。

代码阅读约定

下文代码块从固定 revision 的开源实现中摘取决定语义的关键行;为了聚焦 Q/K/V、cache 和状态更新,省略了异常处理、并行通信、shape 检查和部分无关参数。变量名与核心运算保持源码一致,脚注链接到完整文件和行号。

Dense Softmax 家族

Full Attention

对输入 X,每个 head 独立计算:

1
2
3
4
Qh = X WQh,  Kh = X WKh,  Vh = X WVh
Ah = softmax(Qh Kh^T / sqrt(d) + causal_mask)
Oh = Ah Vh
O = concat(O1, ..., OH) WO

Hugging Face 的 Llama eager 实现把投影、score、softmax 和 value 聚合完整暴露出来:[^llama-code]

1
2
3
4
5
6
7
query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)

attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32)
attn_output = torch.matmul(attn_weights, value_states)
![Full Attention QKV 计算逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/982ed5a640579e999f5d9fbb34221b50.png){ width=100% }
Full/MHA:每个 Query head 拥有独立 K/V,并对全部可见 token 计算 dense score;decode cache 随 `T × Hq` 增长。

瓶颈有两层:prefill 需要形成 T × T score;decode 虽然每步只有一个新 Query,却要读取越来越长的 K/V cache,常被显存带宽限制。

MQA

MQA 保留 Hq 个 Query head,但只生成一组 K/V:

1
2
3
Qi = X WQi,  K = X WK,  V = X WV
Ai = softmax(Qi K^T / sqrt(d))
Oi = Ai V

Llama 的同一个实现通过 num_key_value_heads 表达 MHA、GQA 与 MQA。MQA 的配置边界是 num_key_value_heads = 1repeat_kv 只创建与 Query head 对齐的计算视图,不代表 cache 中保存了 Hq 份 K/V:[^llama-code]

1
2
3
4
5
self.num_key_value_groups = (
config.num_attention_heads // config.num_key_value_heads
)
key_states = repeat_kv(key, module.num_key_value_groups)
value_states = repeat_kv(value, module.num_key_value_groups)
![MQA QKV 计算逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/0f6ed8603111b9b9df54e528bc64826b.png){ width=100% }
MQA:多个 Query head 对同一组 K/V 做 attention。逻辑 score 数量仍由 `Hq` 决定,但 KV cache 从 `Hq` 组降为 1 组。

MQA 的优势集中在 decode cache 和内存读取。代价是所有 Query head 只能从同一份 K/V 表示读取信息,可能损失 head 多样性。

GQA

GQA 在 MHA 与 MQA 之间引入 Hkv 组 K/V,每组服务 Hq / Hkv 个 Query head:

1
2
3
group(i) = floor(i / (Hq / Hkv))
Ai = softmax(Qi Kgroup(i)^T / sqrt(d))
Oi = Ai Vgroup(i)

对应的开源边界仍是 Llama 的两行投影和 repeat_kv:[^llama-code]

1
2
3
4
5
self.k_proj = nn.Linear(hidden_size, num_key_value_heads * head_dim)
self.v_proj = nn.Linear(hidden_size, num_key_value_heads * head_dim)

# Hkv = Hq: MHA; 1 < Hkv < Hq: GQA; Hkv = 1: MQA
key_states = repeat_kv(key, num_attention_heads // num_key_value_heads)
![GQA QKV 计算逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/64a40dcd42043ee11b45cd441c7f5584.png){ width=100% }
GQA:Query head 在组内共享 K/V。`Hkv` 是连续的工程旋钮,而不是与 MHA、MQA 完全分离的算子。

因此 GQA 的准确理解是:用一部分 head 多样性换取一部分 KV cache 和带宽收益。它仍然是对全部历史 token 的 dense attention。

MLA

MLA 不再以“减少 KV head 数”压缩 cache,而是把所有 head 的 K/V 联合投影到低维 latent:

1
2
3
4
cKV = WDKV h
kR = RoPE(WKR h)
kC, vC = WUKV cKV
score = qC kC^T + qR kR^T

DeepSeek-V2 的兼容实现直接展示了联合降投影和位置分量拆分:[^deepseek-v2-code]

1
2
3
4
5
6
7
8
compressed_kv = self.kv_a_proj_with_mqa(hidden_states)
compressed_kv, k_pe = torch.split(
compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1
)
kv = self.kv_b_proj(self.kv_a_layernorm(compressed_kv))
k_nope, value_states = torch.split(
kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1
)

真正体现 memory-optimal decode 的是 DeepSeek 原生推理代码:cache 只分配 kv_lora_rankqk_rope_head_dim;decode 时把 WUKV 的 K 部分吸收到 Query,把 V 部分吸收到输出路径:[^deepseek-v32-mla-code]

1
2
3
4
5
6
7
8
9
self.kv_cache = torch.zeros(max_batch, max_seq, self.kv_lora_rank)
self.pe_cache = torch.zeros(max_batch, max_seq, self.qk_rope_head_dim)

kv, k_pe = torch.split(self.wkv_a(x), [self.kv_lora_rank, self.qk_rope_head_dim], -1)
self.kv_cache[:bsz, start_pos:end_pos] = kv
self.pe_cache[:bsz, start_pos:end_pos] = k_pe.squeeze(2)

q_nope = torch.einsum("bshd,hdc->bshc", q_nope, wkv_b[:, :self.qk_nope_head_dim])
scores = torch.einsum("bshc,btc->bsht", q_nope, self.kv_cache[:bsz, :end_pos])
![MLA QKV 计算逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/6dab7d66272ce9abf28c9e17262ae1ce.png){ width=100% }
MLA:训练或 prefill 可语义上重建各 head 的 K/V;memory-optimal decode 通过矩阵吸收直接在 `cKV` 上计算,并缓存独立的 RoPE 分量 `kR`。

不要从兼容实现误判 cache

某些通用 Transformers 代码会先用 WUKV 展开完整 K/V,再交给标准 attention/cache 接口。这便于兼容,但会掩盖 MLA 的核心收益。判断实际 cache 必须看部署路径中保存的张量,而不只看数学等价的 forward。

MLA 压缩的是 每个 token 的历史表示,并没有选择更少的历史 token。没有 DSA 时,score 仍覆盖全部可见位置。

递归状态家族

GDN

GDN 不保存逐 token 的 K/V。它把历史压进固定矩阵状态,并在每个新 token 到来时执行“先整体衰减,再做一次 delta 写入,最后由 Q 读取”。下面先取参考递推中 H_v=H 的常用情形;新版 FLA 还支持 grouped value heads,此时状态 head 数改为 H_v,递推本身不变。

对象 源码名 Shape 来源与去向 生命周期
当前层输入 hidden_states [B,T,D] 上一子层 → 所有投影 临时激活
归一化 Q/K q, k [B,T,H,K] WQ/WK → ShortConv → SiLU → L2Norm 临时激活
Value v [B,T,H,V] WV → ShortConv → SiLU 临时激活
原始衰减、写入 logits a_proj(h), b_proj(h) [B,T,H] 线性投影 → gate / sigmoid 临时激活
对数衰减、衰减因子、写入率 g、语义 lambda=exp(g)beta [B,T,H] gate 参数化 → 状态更新;lambda 可融合 临时激活 / kernel 中间量
卷积状态 conv_state_q/k/v conv_size 有关的三个短窗口 ShortConv 上一步 → 下一步 有界 cache
递归状态 recurrent_state / h 语义 [B,H,K,V];FLA cache 为 [B,H,V,K] 上一步状态 → 衰减、写入、Q 读取 跨 token 状态
输出 gate、读出与层输出 g_proj(h) / zo / 返回值 [B,T,H,V][B,T,H,V][B,T,D] 状态读取 → Gated RMSNorm → WO 临时激活

a_t=W^a h_tb_t=W^β h_t。FLA 把 Mamba2 风格的衰减参数化融合在 kernel 内:A_logdt_bias 是训练参数,g_t 是非正的对数衰减,λ_t=exp(g_t)∈(0,1]。以下采用默认 allow_neg_eigval=False,因此 beta_t=sigmoid(b_t);开启该选项时实现改为 2 sigmoid(b_t)。把 Q/K 的 L2 归一化和 Query scale 一并写出后,单 token 的物理语义是:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
q_t = L2Norm(SiLU(ShortConv(WQ h)_t)) / sqrt(K)       [B,H,K]
k_t = L2Norm(SiLU(ShortConv(WK h)_t)) [B,H,K]
v_t = SiLU(ShortConv(WV h)_t) [B,H,V]

g_t = -exp(A_log) * softplus(Wa h_t + dt_bias) [B,H]
lambda_t = exp(g_t) [B,H]
beta_t = sigmoid(Wbeta h_t) [B,H]
z_t = Wg h_t [B,H,V]

S_bar_t = lambda_t S_(t-1) [B,H,K,V]
v_hat_t = S_bar_t^T k_t [B,H,V]
e_t = beta_t (v_t - v_hat_t) [B,H,V]
S_t = S_bar_t + k_t e_t^T [B,H,K,V]
r_t = q_t^T S_t [B,H,V]
y_t = WO(GatedRMSNorm(r_t, z_t)) [B,D]

其中 v_hat_t 是旧状态对当前 key 的预测,e_t 是要纠正的 value 残差,k_t e_t^T 是秩 1 写入。**S_bar_tv_hat_te_t 在 naive 代码里是可观察中间量;优化 kernel 可以融合它们,但不能从计算语义中省略。**

FLA 层先物化投影与三个短卷积输出,再把原始 gate、训练参数和旧状态交给 fused/chunk recurrence:[^gdn-layer-code]

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
q, conv_state_q = self.q_conv1d(
x=self.q_proj(hidden_states), cache=conv_state_q,
output_final_state=use_cache, cu_seqlens=cu_seqlens,
)
k, conv_state_k = self.k_conv1d(
x=self.k_proj(hidden_states), cache=conv_state_k,
output_final_state=use_cache, cu_seqlens=cu_seqlens,
)
v, conv_state_v = self.v_conv1d(
x=self.v_proj(hidden_states), cache=conv_state_v,
output_final_state=use_cache, cu_seqlens=cu_seqlens,
)
beta = self.b_proj(hidden_states)
o, recurrent_state = chunk_gated_delta_rule(
q=q, k=k, v=v, g=self.a_proj(hidden_states), beta=beta,
A_log=self.A_log, dt_bias=self.dt_bias,
initial_state=recurrent_state, output_final_state=use_cache,
use_qk_l2norm_in_kernel=True, use_gate_in_kernel=True,
use_beta_sigmoid_in_kernel=True,
allow_neg_eigval=self.allow_neg_eigval,
state_v_first=True, cu_seqlens=cu_seqlens,
)

state_v_first=True 只把 cache 的物理布局从语义上的 [K,V] 转为 [V,K],公式仍按 [K,V] 书写;cu_seqlens 仅描述 packed sequence 边界。短序列推理会切到 fused_recurrent_gated_delta_rule,参数与状态对象相同。

上面的 chunk_gated_delta_rule 是训练友好的并行实现;下列 naive 循环才直接对应前述对象。代码中的 h 就是公式的 Sb_v 先表示 v_t,随后被原地改写为 e_t:[^gdn-naive-code]

1
2
3
4
5
6
7
h = h.clone() * g[:, :, i].exp()[..., None, None]
b_v = v[:, :, i].clone()
b_k = k[:, :, i]
b_v = b_v - (h.clone() * b_k[..., None]).sum(-2)
b_v = b_v * beta[:, :, i][..., None]
h = h.clone() + b_k.unsqueeze(-1) * b_v.unsqueeze(-2)
o[:, :, i] = torch.einsum("bhd,bhdm->bhm", q[:, :, i], h)
![GDN QKV 与状态更新逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/dc34779e03cd2aef72eb515bde953376.png){ width=100% }
自绘物理语义图,依据 FLA commit `c70f11c5`:圆点是张量,双圈是跨 token 状态,方块是实际算子,权重和超参标在线上。GDN 的 `lambda_t[B,H]` 对每个 head 的整张 `K×V` 状态统一衰减。

固定状态不是无损 KV 压缩。多个历史 token 会竞争同一矩阵容量;除 S_t 外,部署还要保存三个短卷积的有界状态。训练时 chunk kernel 并行计算多个时间步,decode 时 recurrent kernel 逐 token 更新,但两条路径都实现同一递推。

KDA

KDA 保留 GDN 的 delta 写入,却把“每个 head 一个衰减标量”改成“每个 value head、每个 key channel 一个衰减值”。令 Q/K head 数为 H,Value/状态 head 数为 H_v,分组倍率 G=H_v/H;Kimi Linear 模型取 G=1,FLA naive 实现允许 G>1 并把 Q/K 复制到对应的 value heads。

对象 源码名 Shape 来源与去向 生命周期
Q/K q, k [B,T,H,K] 线性投影 → ShortConv/SiLU → L2Norm 临时激活
展开后的 Q/K repeat_interleave(q/k, G) [B,T,H_v,K] head 对齐 → 状态读写 naive 中物化;fused kernel 可隐式
Value v [B,T,H_v,V] WV → ShortConv → SiLU 临时激活
原始 channel gate f_b_proj(f_a_proj(h)) [B,T,H_v K] 两层低秩投影 → fused_kda_gate 临时激活
对数衰减、衰减因子 g、语义 alpha=exp(g) [B,T,H_v,K] A_log/dt_bias 参数化 → 对角衰减;alpha 可融合 临时激活 / kernel 中间量
写入率 beta [B,T,H_v] sigmoid(Wbeta h) → delta 残差 临时激活
卷积与递归状态 conv_states, recurrent_state 三个短窗口;[B,H_v,K,V] 上一步 → 当前步 → 下一步 有界 cache / 跨 token 状态
输出 gate 与层输出 g_a/g_bo [B,T,H_v,V][B,T,D] 状态读出 → Gated RMSNorm → WO 临时激活

这里有两个容易混淆的 gf_a/f_b 生成的是 state decay gateg_a/g_b 生成的是 output gate,二者参数和用途完全不同。写成完整递推:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
q_t  = L2Norm(SiLU(ShortConv(WQ h)_t))                 [B,H,K]
k_t = L2Norm(SiLU(ShortConv(WK h)_t)) [B,H,K]
v_t = SiLU(ShortConv(WV h)_t) [B,Hv,V]
q'_t = repeat(q_t, G) / sqrt(K) [B,Hv,K]
k'_t = repeat(k_t, G) [B,Hv,K]

g_t = -exp(A_log) * softplus(Wfb Wfa h_t + dt_bias) [B,Hv,K]
alpha_t = exp(g_t) [B,Hv,K]
beta_t = sigmoid(Wbeta h_t) [B,Hv]
z_t = Wgb Wga h_t [B,Hv,V]

S_bar_t = Diag(alpha_t) S_(t-1) [B,Hv,K,V]
v_hat_t = S_bar_t^T k'_t [B,Hv,V]
e_t = beta_t (v_t - v_hat_t) [B,Hv,V]
S_t = S_bar_t + k'_t e_t^T [B,Hv,K,V]
r_t = q'_t^T S_t [B,Hv,V]
y_t = WO(GatedRMSNorm(r_t, z_t)) [B,D]

Kimi Linear 的模型代码把 Q/K/V 短卷积、channel gate、写入率、递归状态和输出 gate 全部暴露出来:[^kimi-model-code]

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
q, conv_state_q = self.q_conv1d(
x=self.q_proj(hidden_states), cache=conv_state_q,
output_final_state=use_cache, cu_seqlens=cu_seqlens,
)
k, conv_state_k = self.k_conv1d(
x=self.k_proj(hidden_states), cache=conv_state_k,
output_final_state=use_cache, cu_seqlens=cu_seqlens,
)
v, conv_state_v = self.v_conv1d(
x=self.v_proj(hidden_states), cache=conv_state_v,
output_final_state=use_cache, cu_seqlens=cu_seqlens,
)
g = self.f_b_proj(self.f_a_proj(hidden_states))
g = fused_kda_gate(g, self.A_log, self.head_dim, g_bias=self.dt_bias)
beta = self.b_proj(hidden_states).float().sigmoid()
o, recurrent_state = chunk_kda(
q=q, k=k, v=v, g=g, beta=beta,
initial_state=recurrent_state, output_final_state=True,
use_qk_l2norm_in_kernel=True, cu_seqlens=cu_seqlens,
)
g = self.g_b_proj(self.g_a_proj(hidden_states))
g = rearrange(g, "... (h d) -> ... h d", d=self.head_dim)
o = self.o_norm(o, g)
o = self.o_proj(rearrange(o, "b t h d -> b t (h d)"))

最后三行是源码对变量名 g 的复用:此时它已经不再是 state decay gate,而是 output gate。短序列 decode 使用 fused_recurrent_kda,与示出的 chunk_kda 接收相同的 q/k/v/g/beta/initial_state,只改变调度与物化方式。

FLA naive KDA 进一步说明了 G、对角衰减和实际收缩轴。(k_i[...,None] * S).sum(-2) 沿 K 维收缩得到 [B,H_v,V],外积再恢复 [B,H_v,K,V]:[^kda-naive-code]

1
2
3
4
5
6
7
8
9
10
11
G = HV // H
q = q.repeat_interleave(G, dim=2) * scale
k = k.repeat_interleave(G, dim=2)
S = k.new_zeros(B, HV, K, V)
for i in range(T):
q_i, k_i, v_i = q[:, i], k[:, i], v[:, i]
g_i, b_i = g[:, i], beta[:, i]
S = S * g_i[..., None].exp()
residual = v_i - (k_i[..., None] * S).sum(-2)
S = S + torch.einsum("bhk,bhv->bhkv", b_i[..., None] * k_i, residual)
o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S)
![KDA QKV 与状态更新逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/e49e46af37d99508b050dad51534ed0d.png){ width=100% }
自绘物理语义图,依据 Kimi Linear revision `e1df551a` 与 FLA commit `c70f11c5`:KDA 的 `alpha_t[B,Hv,K]` 沿状态矩阵的 K 轴逐 channel 衰减;`beta_t[B,Hv]` 仍控制整个 value 残差的写入强度。

因此 KDA 与 GDN 的差异不在“是否有 Q/K/V”,而在状态转移矩阵:GDN 是 lambda_t I,KDA 是 Diag(alpha_t)。更细粒度 gate 提高了记忆方向的独立性,也增加 gate 生成和状态带宽;Kimi Linear 因而采用 3 层 KDA 配 1 层 MLA 的混合结构,而不是让固定状态完全替代精确 token 检索。

Token 稀疏家族

DSA

DSA 不是另一套 K/V 压缩,而是在 MLA 的两份 cache 旁增加一套轻量索引 cache。Lightning Indexer 先对全部历史位置打低维分数,再把 top-k 位置交给昂贵的主 MLA。

对象 源码名 Shape(单卡逻辑) 来源与去向 生命周期
当前输入、Q latent x, qr [B,S,D][B,S,Rq] Wq_a → RMSNorm,同时供 MLA Q 与 Indexer Q 临时激活
主 MLA Query q_nope, q_pe [B,S,Hm,Dc][B,S,Hm,Dr] Wq_b → Split → RoPE 临时激活
MLA 压缩 KV、位置 K kv, k_pe [B,S,Rkv][B,S,Dr] Wkv_a → Split/RMSNorm/RoPE 写入 MLA cache
Indexer Query q, q_fp8, q_scale [B,S,Hi,Di] 及其 FP8/scale WIQ qr → Split/RoPE/Concat → Hadamard → Quantize 临时激活
Indexer Key k, k_fp8, k_scale [B,S,Di] 及其 FP8/scale WIK x → Norm/Split/RoPE/Concat → Hadamard → Quantize 写入 Indexer cache
head 权重 weights [B,S,Hi] WIW x / sqrt(Hi) → FP8 index kernel 临时激活
索引分数、位置 index_score, topk_indices [B,S,T][B,S,k] 全历史 index score → TopK 临时激活 / 离散索引
主 MLA 概率与 latent 聚合 scores, x [B,S,Hm,T][B,S,Hm,Rkv] 稀疏 mask/Softmax → 压缩 cache 聚合 临时激活
层输出 返回 x [B,S,D] WUV 解码 value → WO 临时激活

其中 H_iD_i 是 Indexer head 数和 head 维度;H_m 是 MLA head 数;R_q/R_kv 是 Q/KV latent rank;D_c/D_r 是 MLA 的 NoPE/RoPE Query 维度。Indexer 的论文公式包含一个不可省略的 ReLU:[^dsa-release]

1
2
3
4
5
6
qI_(t,h) = Hadamard(Concat(RoPE(qI_R), qI_C))          [B,Hi,Di]
kI_j = Hadamard(Concat(RoPE(kI_R), kI_C)) [B,Di]
w_(t,h) = (WIW h_t)_h / sqrt(Hi) [B,Hi]

I_(t,j) = sum_(h=1..Hi) w_(t,h) ReLU(<qI_(t,h), kI_j> / sqrt(Di))
J_t = TopK_j(I_(t,j), k) [B,k]

开源 kernel 先执行 FP8 点积,再做逐 head ReLU、乘权重、跨 head 求和,最后乘 Indexer K 的量化 scale;模型侧的 weights 已合入 Query scale 与 D_i^-0.5:[^deepseek-v32-kernel-code]

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
q = self.wq_b(qr).view(bsz, seqlen, self.n_heads, self.head_dim)
q_pe, q_nope = torch.split(q, [self.rope_head_dim, self.head_dim - self.rope_head_dim], dim=-1)
q_pe = apply_rotary_emb(q_pe, freqs_cis, False)
q = rotate_activation(torch.cat([q_pe, q_nope], dim=-1))

k = self.k_norm(self.wk(x))
k_pe, k_nope = torch.split(k, [self.rope_head_dim, self.head_dim - self.rope_head_dim], dim=-1)
k_pe = apply_rotary_emb(k_pe.unsqueeze(2), freqs_cis, False).squeeze(2)
k = rotate_activation(torch.cat([k_pe, k_nope], dim=-1))

q_fp8, q_scale = act_quant(q, block_size, self.scale_fmt)
k_fp8, k_scale = act_quant(k, block_size, self.scale_fmt)
self.k_cache[:bsz, start_pos:end_pos] = k_fp8
self.k_scale_cache[:bsz, start_pos:end_pos] = k_scale
weights = self.weights_proj(x.float()) * self.n_heads ** -0.5
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
index_score = fp8_index(q_fp8, weights, self.k_cache[:bsz, :end_pos], self.k_scale_cache[:bsz, :end_pos])
topk_indices = index_score.topk(min(self.index_topk, end_pos), dim=-1)[1]

top-k 随后约束主 MLA。对 memory-optimal decode,可把 WUK 吸收到 q_nope,把 WUV 放到 latent 聚合之后:

1
2
3
4
5
qC'_t = qC_t WUK                                       [B,Hm,Rkv]
a_(t,j) = (qC'_t cKV_j^T + qR_t kR_j^T) * scale [B,Hm,T]
p_t = softmax(a_t + mask(J_t)) [B,Hm,T]
u_t = sum_j p_(t,j) cKV_j [B,Hm,Rkv]
y_t = WO(Concat_h(u_t WUV_h)) [B,D]

但必须区分 算法目标这份 Python 参考代码实际物化的路径。论文目标是每个 Query 只 gather J_t 中的 k 个 latent,再计算 a/u,把整段主注意力从 O(ST) 降为 O(Sk),单 token decode 则从 O(T) 降为 O(k);固定 revision 的开源 Python 代码却先对全部 T 个位置形成 dense scores,再 scatter 稀疏 mask。它验证选择语义,但本身不是获得这部分稀疏收益的 production kernel:[^deepseek-v32-mla-code]

1
2
3
4
5
6
7
8
topk_indices = self.indexer(x, qr, start_pos, freqs_cis, mask)
index_mask = torch.full((bsz, 1, end_pos), float("-inf"), device=x.device)
index_mask = index_mask.scatter_(-1, topk_indices, 0)
scores += index_mask.unsqueeze(2)
scores = scores.softmax(dim=-1)
x = torch.einsum("bsht,btc->bshc", scores, self.kv_cache[:, :end_pos])
x = torch.einsum("bshc,hdc->bshd", x, wkv_b[:, -self.v_head_dim:])
x = self.wo(x.flatten(2))
![DSA Indexer 与 MLA 计算逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/c37f03a421280a9571d600b223e45b10.png){ width=100% }
自绘 decode 对象图,依据 DeepSeek-V3.2-Exp commit `87e509a2`:Indexer 物理缓存 FP8 K 与 scale,MLA 物理缓存 `cKV/kR`。实线表示开源 Python 的 dense-score+mask 路径;虚线 gather 表示论文目标的 production sparse kernel 等价路径。

稀疏的边界

DSA 的算法收益来自 主 MLA 的高维 score/value 只读取 top-k;Indexer 仍扫描全部历史。若实现像上述 Python reference 一样先算 dense 主 score 再 mask,语义仍是 DSA,但不会得到完整的稀疏 kernel 省算。index_topk=2048 是 DeepSeek-V3.2-Exp 的具体配置,不是 DSA 定义中的固定常数。

混合架构

长上下文模型越来越少押注单一 Attention。线性或递归状态层负责大部分低成本历史传播,周期性 Full/MLA 层补充精确 token 检索,是更现实的质量—效率折中。

![Kimi Linear 混合架构](https://pic.shaojiemike.top/shaojiemike/2026/07/0be2c372768092289c465956c0b05bdf.png){ width=76% }
Kimi Linear 论文 Figure 3:以三层 KDA 配一层 MLA 的混合结构。它说明 KDA 的目标是承担大部分历史传播,而非证明 softmax attention 已无必要。
![Kimi Linear 长上下文效率](https://pic.shaojiemike.top/shaojiemike/2026/07/baa3d8e8bdd4dc6e8a5016537f6fea88.png){ width=86% }
Kimi Linear 论文 Figure 7:图中 1M context 的可见标注约为 prefill 2.9×、decode TPOT 2.2×(相对 MLA)。论文摘要/正文另有“最高 6× decode”的表述,两者属于不同统计或边界,不能直接合并。

这种混合思想也适用于其他组合:GQA 降低 cache,DSA 降低主 attention 的 token 数,二者解决不同维度;MLA 与 DSA 更是明确的基座和索引关系。

选择框架

主要瓶颈 优先考察 原因 必验风险
Decode KV cache / 带宽 GQA、MQA、MLA 直接减少每 token 历史表示 质量回退、kernel 与 cache layout
极长上下文的固定内存 GDN、KDA recurrent state 不随 T 增长 状态容量、训练 kernel、跨序列边界
极长上下文主 attention FLOPs DSA 主 MLA 只读取 top-k Indexer recall、top-k 通信、稀疏 kernel
质量优先、上下文适中 Full/MHA、GQA 精确 dense retrieval,生态成熟 cache 和带宽上限
质量与极长上下文同时要求 Hybrid 递归层传播、周期性 softmax 检索 层型比例、状态/cache 双系统

工程选型时至少测四件事:

  1. Prefill 与 decode 分开:二者的算力、带宽和并行特性完全不同。
  2. 峰值 cache 与实际吞吐同时测:更小的理论 cache 可能被不合适的 layout、dequant 或 gather 抵消。
  3. 固定准确率条件下比较:MQA/GQA group 数、MLA rank、DSA top-k、GDN/KDA state 维度都属于质量—成本旋钮。
  4. 按实现边界比较:FlashAttention、FlashMLA、chunk recurrence 与稀疏 gather 的 kernel 成熟度往往比大 O 公式更决定真实速度。

常见误区

  • “MQA/GQA 是稀疏 attention”:错误。它们共享 K/V head,但仍读取所有历史 token。
  • “MLA 把 attention 变成线性复杂度”:错误。MLA 压缩每个 token 的 K/V 表示,没有压缩 token 数。
  • “GDN/KDA 只是更激进的 KV cache 压缩”:不准确。它们改变了历史记忆算法,用固定状态递推替代逐 token softmax 检索。
  • “DSA 只算 top-k,所以模块没有全历史成本”:不完整。主 MLA 是 top-k,Indexer 仍扫描历史。
  • “后出现的结构必然取代前者”:错误。真实模型广泛使用混合层,不同方法解决的是 cache、状态容量、token 连通性或 kernel 效率中的不同问题。

总结

Attention 的发展路径可以压缩成一句话:先共享 K/V,再联合压缩 K/V;另一条路线把历史写入固定状态;最新的稀疏路线先检索,再让昂贵的主注意力读取少数 token。

理解这些结构时,不要只背缩写。顺着代码追四个对象就足够:

  1. Q/K/V 如何投影、按多少 head 组织;
  2. cache 保存完整 K/V、latent、Indexer K,还是固定状态;
  3. score 是对全部历史、top-k,还是根本不存在显式 T×T 矩阵;
  4. 被省掉的成本是否转移到了低秩重建、状态更新、索引器或稀疏 gather。

本文七张计算图均依据版本固定源码自绘,SVG 源文件与 1920 像素宽 PNG 同时保留,并通过结构、箭头、碰撞、语义几何与版式校验。

参考文献

[^mha-paper]: Vaswani et al., Attention Is All You Need, first submitted 2017-06-12.
[^mqa-paper]: Shazeer, Fast Transformer Decoding: One Write-Head is All You Need, first submitted 2019-11-06.
[^gqa-paper]: Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints, first submitted 2023-05-22.
[^mla-paper]: DeepSeek-AI, DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model, first submitted 2024-05-07.
[^gdn-paper]: Yang et al., Gated Delta Networks: Improving Mamba2 with Delta Rule, first submitted 2024-12-09.
[^dsa-release]: DeepSeek-AI, DeepSeek-V3.2-Exp release, 2025-09-29; model and paper repository.
[^kda-paper]: Zhang et al., Kimi Linear: An Expressive, Efficient Attention Architecture, first submitted 2025-10-30.
[^llama-code]: Hugging Face Transformers, modeling_llama.py, commit 29985e67.
[^deepseek-v2-code]: DeepSeek-AI, modeling_deepseek.py, revision 604d5664.
[^deepseek-v32-mla-code]: DeepSeek-AI, inference/model.py MLA, commit 87e509a2.
[^deepseek-v32-kernel-code]: DeepSeek-AI, inference/kernel.py FP8 Indexer, commit 87e509a2.
[^gdn-layer-code]: FLA, fla/layers/gated_deltanet.py, commit c70f11c5.
[^gdn-naive-code]: FLA, fla/ops/gated_delta_rule/naive.py, commit c70f11c5.
[^kimi-model-code]: Moonshot AI, modeling_kimi.py, revision e1df551a.
[^kda-naive-code]: FLA, fla/ops/kda/naive.py, commit c70f11c5.

Author

Shaojie Tan

Posted on

2026-07-21

Updated on

2026-07-21

Licensed under